Skip to main content

Auto-Encoding Variational Bayes

路线图​

想让模型分布接近真实数据分布
↓
最大化训练图片在模型下的对数似然
↓
用“先采样潜变量 z,再生成图片 x”的方式建模
↓
计算 pθ(x) 时必须对 z 积分,难以直接优化
↓
直接从先验 p(z) 蒙特卡洛可估计 pθ(x),但有限采样下 log 估计有偏且常与当前 x 错配
↓
引入编码器 qφ(z|x) 近似不可得的真实后验
↓
得到 log pθ(x) 的可优化下界 ELBO
↓
重建期望从 qφ(z|x) 采样估计,并用重参数化回传梯度;KL 项有闭式解
↓
最小化负 ELBO(MSE/BCE 重建项 + KL 项)的 batch 平均
↓
训练后从先验 p(z) 采样,再由解码器生成新图片

0. 记号:一张图、一个 batch、多个 zz​

  • 训练集:D={x(n)}n=1N\mathcal D=\{x^{(n)}\}_{n=1}^{N},其中 x(n)x^{(n)} 是第 nn 张真实训练图片。
  • xx:后续推导中,为简洁起见,固定的一张训练图片 x=x(n)x=x^{(n)};它不是模型输出,也不是整个 batch。
  • x^\hat{x}:解码器针对 zz 生成的重建结果(或输出分布的均值/参数)。
  • zz:连续潜变量,表示影响图像的隐藏因素,如姿态、光照、物体属性或高层语义。
  • p(z)p(z):潜变量先验,通常设为 N(0,I)\mathcal N(0,I)。
  • pθ(x∣z)p_\theta(x\mid z):解码器/生成模型;给定 zz 后生成图片 xx 的条件分布。
  • qϕ(z∣x)q_\phi(z\mid x):编码器/近似后验;给定图片 xx 后预测其潜变量的分布。

实际训练有两个独立的数量:

符号表示什么常见取值
BB一个 batch 的图片数:x(1),…,x(B)x^{(1)},\ldots,x^{(B)}32、64、128 等
KK对每一张图片采样的潜变量数通常为 1

完整的采样记号是:

z(b,k)∼qϕ(z∣x(b)),b=1,…,B,k=1,…,K.z^{(b,k)}\sim q_\phi(z\mid x^{(b)}), \qquad b=1,\ldots,B,\quad k=1,\ldots,K.

下文先对单张固定图片 xx 推导,最后再回到 batch。

1. 为什么目标是最大化数据似然​

对于图片空间 我们可以分为两部分 一种是我们希望能从模型中预测出来的图片,他们应该是 合理的/有意义的/特定领域 图片,而不是 随机/非自然/领域外 的像素组合

从训练数据(合理的/有意义的/特定领域的)得到的数据分布 pdata(x)p_{\mathrm{data}}(x);生成模型定义一个可学习的分布 pθ(x)p_\theta(x)。我们希望模型分布接近数据分布,一种具有明确统计意义的选择是最小化:

DKL ⁣(pdata ∥ pθ)=Epdata[log⁡pdata(x)]⏟不随 θ 变化−Epdata[log⁡pθ(x)].\begin{aligned} D_{\mathrm{KL}}\!\left(p_{\mathrm{data}}\,\|\,p_\theta\right) &=\underbrace{\mathbb E_{p_{\mathrm{data}}} \left[\log p_{\mathrm{data}}(x)\right]}_{\text{不随 }\theta\text{ 变化}} -\mathbb E_{p_{\mathrm{data}}}\left[\log p_\theta(x)\right]. \end{aligned}

第一项与参数 θ\theta 无关,因此最小化该 KL 等价于最大化:

max⁡θ  Ex∼pdata[log⁡pθ(x)].\max_\theta\;\mathbb E_{x\sim p_{\mathrm{data}}}[\log p_\theta(x)].

我们不知道 pdatap_{\mathrm{data}},但有从中采样得到的训练集,因此用样本平均近似:

max⁡θ  1N∑n=1Nlog⁡pθ(x(n)).\max_\theta\;\frac1N\sum_{n=1}^{N}\log p_\theta\bigl(x^{(n)}\bigr).

这就是最大似然的动机:让模型分布接近数据分布。

2. 用潜变量建模图片​

直接描述复杂图片分布 pθ(x)p_\theta(x) 很困难。VAE 采用生成假设:

先从简单分布中抽取隐藏因素 zz,再根据 zz 生成图片 xx。

z∼p(z)⟶x∼pθ(x∣z).z\sim p(z) \quad\longrightarrow\quad x\sim p_\theta(x\mid z).

它给出联合分布:

pθ(x,z)=pθ(x∣z)p(z).p_\theta(x,z)=p_\theta(x\mid z)p(z).

其中 p(z)p(z) 是我们预先指定、且容易采样的潜变量分布,通常取标准高斯 N(0,I)\mathcal N(0,I)。它不是由某张图片 xx 决定的:生成新图片时,先从这个固定分布抽取 zz,再送入解码器。选择简单先验的目的,是让“生成时该从哪里抽取 zz”成为一个明确且容易执行的问题。

图片本身的概率必须将潜变量边缘化:

pθ(x)=∫pθ(x,z) dz=∫pθ(x∣z)p(z) dz.p_\theta(x)=\int p_\theta(x,z)\,dz =\int p_\theta(x\mid z)p(z)\,dz.

3. 难点:边缘似然与真实后验都难算​

对神经网络解码器,上述积分通常没有闭式解,因而无法直接精确优化 log⁡pθ(x)\log p_\theta(x)。

从先验 z(k)∼p(z)z^{(k)}\sim p(z) 做 Monte Carlo 可以估计概率本身:

pθ(x)=Ep(z)[pθ(x∣z)]≈1K∑k=1Kpθ(x∣z(k)).p_\theta(x) =\mathbb E_{p(z)}[p_\theta(x\mid z)] \approx\frac1K\sum_{k=1}^{K}p_\theta\bigl(x\mid z^{(k)}\bigr).

但训练目标是 log⁡pθ(x)\log p_\theta(x)。对样本均值取对数不是 log⁡pθ(x)\log p_\theta(x) 的无偏估计;由 Jensen 不等式

≤log⁡E[p^K(x)]=log⁡pθ(x).\leq \log\mathbb E[\hat p_K(x)] = \log p_\theta(x).

它在有限 KK 时实际上给出了一个偏低的下界。更关键的是先验 p(z)p(z) 不知道当前图片 xx,大多数先验样本无法解释它,只有极少数样本有显著权重。

核心澄清:问题不只是“有偏与无偏”

直接先验采样要处理的是:

log⁡Ep(z)[pθ(x∣z)].\log\mathbb E_{p(z)}[p_\theta(x\mid z)].

先用样本均值估计内部期望,再取对数,得到的 log⁡p^θ(x)\log\hat p_\theta(x) 在有限 KK 下是对 log⁡pθ(x)\log p_\theta(x) 的向下有偏估计。它并非理论上不能使用:K→∞K\to\infty 时会收敛;真正的实践困难还包括先验与后验的错配——对固定 xx,能解释它的区域 pθ(z∣x)p_\theta(z\mid x) 往往只占 p(z)p(z) 的很小部分,因此有限预算下可能几乎抽不到有效样本。

而 VAE 重建项处理的是另一个目标:

Eqϕ(z∣x)[log⁡pθ(x∣z)].\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)].

对这个期望直接取样本均值是无偏的。但这不表示 VAE 得到了 log⁡pθ(x)\log p_\theta(x) 的无偏估计;VAE 改为优化 ELBO:

LELBO(x)≤log⁡pθ(x).\mathcal L_{\mathrm{ELBO}}(x)\leq\log p_\theta(x).

ELBO 对真实对数似然而言是下界(偏低),但其组成部分能稳定地计算或估计。qϕ(z∣x)q_\phi(z\mid x) 的作用是使有限采样更有效;KL 项则约束它不能只为重建而任意偏离先验。

我们真正想知道的是:给定图片 xx,哪些 zz 更可能生成它?这对应真实后验:

pθ(z∣x)=pθ(x,z)pθ(x)=pθ(x∣z)p(z)∫pθ(x∣z)p(z) dz.p_\theta(z\mid x) =\frac{p_\theta(x,z)}{p_\theta(x)} =\frac{p_\theta(x\mid z)p(z)} {\int p_\theta(x\mid z)p(z)\,dz}.

其分母正是难算的 pθ(x)p_\theta(x),所以真实后验也不可直接得到。

4. 引入编码器 qϕ(z∣x)q_\phi(z\mid x)​

VAE 用一个神经网络 qϕ(z∣x)q_\phi(z\mid x) 近似真实后验 pθ(z∣x)p_\theta(z\mid x)。常用对角高斯:

qϕ(z∣x)=N(μϕ(x),diag⁡(σϕ2(x))).q_\phi(z\mid x)= \mathcal N\bigl(\mu_\phi(x),\operatorname{diag}(\sigma_\phi^2(x))\bigr).

编码器输入当前图片 xx,输出 μϕ(x)\mu_\phi(x) 与 log⁡σϕ2(x)\log\sigma_\phi^2(x)(logvar)。和不看 xx 的先验相比,qϕ(z∣x)q_\phi(z\mid x) 会把采样集中到更可能解释当前图片的区域。

5. 从后验近似得到 ELBO​

考虑近似后验与真实后验的 KL 散度:

DKL(qϕ(z∣x)∥pθ(z∣x))=Eqϕ(z∣x)[log⁡qϕ(z∣x)pθ(z∣x)]=Eqϕ(z∣x)[log⁡qϕ(z∣x)pθ(x,z)/pθ(x)]=log⁡pθ(x)−Eqϕ(z∣x)[log⁡pθ(x,z)qϕ(z∣x)].\begin{aligned} D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p_\theta(z\mid x)\bigr) &=\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{q_\phi(z\mid x)}{p_\theta(z\mid x)}\right]\\ &=\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{q_\phi(z\mid x)}{p_\theta(x,z)/p_\theta(x)}\right]\\ &=\log p_\theta(x) -\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{p_\theta(x,z)}{q_\phi(z\mid x)}\right]. \end{aligned}

最后一步中 log⁡pθ(x)\log p_\theta(x) 与积分变量 zz 无关,因此能从期望中移出。定义:

LELBO(x)=Eqϕ(z∣x)[log⁡pθ(x,z)qϕ(z∣x)].\mathcal L_{\mathrm{ELBO}}(x) =\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{p_\theta(x,z)}{q_\phi(z\mid x)}\right].

便有精确恒等式:

log⁡pθ(x)=LELBO(x)⏟证据下界+DKL(qϕ(z∣x)∥pθ(z∣x))⏟≥0.\log p_\theta(x) =\underbrace{\mathcal L_{\mathrm{ELBO}}(x)}_{\text{证据下界}} +\underbrace{D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p_\theta(z\mid x)\bigr)}_{\geq 0}.

因此 LELBO(x)≤log⁡pθ(x)\mathcal L_{\mathrm{ELBO}}(x)\leq\log p_\theta(x)。最大化 ELBO 会提高似然下界;若 qϕ(z∣x)q_\phi(z\mid x) 恰好等于真实后验,二者相等。

6. ELBO 为什么是“重建 − KL”​

代入 pθ(x,z)=pθ(x∣z)p(z)p_\theta(x,z)=p_\theta(x\mid z)p(z):

LELBO(x)=Eqϕ(z∣x)[log⁡pθ(x∣z)p(z)qϕ(z∣x)]=Eqϕ(z∣x)[log⁡pθ(x∣z)]+Eqϕ(z∣x)[log⁡p(z)qϕ(z∣x)]=Eqϕ(z∣x)[log⁡pθ(x∣z)]⏟重建项−DKL(qϕ(z∣x)∥p(z))⏟先验匹配项.\begin{aligned} \mathcal L_{\mathrm{ELBO}}(x) &=\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{p_\theta(x\mid z)p(z)}{q_\phi(z\mid x)}\right]\\ &=\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)] +\mathbb E_{q_\phi(z\mid x)} \left[\log\frac{p(z)}{q_\phi(z\mid x)}\right]\\ &=\underbrace{\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)]}_{\text{重建项}} -\underbrace{D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p(z)\bigr)}_{\text{先验匹配项}}. \end{aligned}
项约束什么为什么需要它
重建项 Eq[log⁡pθ(x∣z)]\mathbb E_q[\log p_\theta(x\mid z)]从 zz 能否还原当前 xx让 zz 保留图片信息
KL 项 DKL(qϕ(z∣x)∥p(z))D_{\mathrm{KL}}(q_\phi(z\mid x)\|p(z))编码出的分布是否接近固定先验保证生成时从 p(z)p(z) 采样仍有效

这两个目标缺一不可:

  • 只有重建项:编码器只需为每张图找到任何能让解码器复原它的 zz。不同图片的编码区域可以相互零散、没有统一形状;解码器只见过这些区域,而生成时却要从 p(z)=N(0,I)p(z)=\mathcal N(0,I) 抽样,因此抽到的 zz 可能落在从未训练过的空洞中。这就是普通 AE 虽能重建、却不一定能随机生成的原因。
  • 只有 KL 项:最容易的解是对所有图片都令 qϕ(z∣x)=p(z)q_\phi(z\mid x)=p(z)。此时编码器输出不再依赖 xx,zz 不携带图片信息;又没有重建项训练解码器保留 xx,模型自然无法完成重建。

先验的选择、网络结构、近似分布是人为设计;选定这套模型和下界后,“重建减 KL”是数学结果。“重建”和“先验匹配”则是我们对这个结果的直观解释。

7. 为什么这个目标可训练​

ELBO 仍含 zz,但它的两部分可以实际计算或估计。

7.1 重建项:从 qϕ(z∣x)q_\phi(z\mid x) 采样​

对同一张固定图片 xx,独立采样 KK 个潜变量:

z(k)∼qϕ(z∣x),k=1,…,K.z^{(k)}\sim q_\phi(z\mid x),\qquad k=1,\ldots,K.

则:

Eqϕ(z∣x)[log⁡pθ(x∣z)]≈1K∑k=1Klog⁡pθ(x∣z(k)).\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)] \approx\frac1K\sum_{k=1}^{K}\log p_\theta\bigl(x\mid z^{(k)}\bigr).

当 z(k)z^{(k)} 独立地从 qϕ(z∣x)q_\phi(z\mid x) 采样时,样本均值是该期望的无偏估计。每个 z(k)z^{(k)} 都参与求和,但不代表它们的数值贡献相同。相比从先验采样,qϕ(z∣x)q_\phi(z\mid x) 的样本与当前 xx 更相关,通常具有更低方差。实践中常取 K=1K=1。

例子:为什么从 p(z)p(z) 直接估计似然很困难,而从 qϕ(z∣x)q_\phi(z\mid x) 估计重建项可行?

先用一个极端但直观的离散例子。令潜变量只能取 100 个值:

z∈{1,2,…,100},p(z=j)=1100.z\in\{1,2,\ldots,100\}, \qquad p(z=j)=\frac{1}{100}.

固定一张训练图片 xx。假设只有 z=17z=17 能生成它:

pθ(x∣z=17)=1,pθ(x∣z≠17)=0.p_\theta(x\mid z=17)=1, \qquad p_\theta(x\mid z\ne17)=0.

那么真实边缘似然是:

pθ(x)=∑zpθ(x∣z)p(z)=1×1100=0.01.p_\theta(x)=\sum_z p_\theta(x\mid z)p(z)=1\times\frac1{100}=0.01.

若直接从先验 p(z)p(z) 抽 10 次,抽到 z=17z=17 的概率只有:

1−(99100)10≈9.6%.1-\left(\frac{99}{100}\right)^{10}\approx9.6\%.

也就是说约 90.4% 的时候,10 个样本都没有任何贡献,得到的 p^θ(x)=0\hat p_\theta(x)=0,再取 log⁡p^θ(x)\log\hat p_\theta(x) 会变成 −∞-\infty;即使真实 log⁡pθ(x)=log⁡0.01\log p_\theta(x)=\log0.01 是有限的。这说明问题不只是“有一点方差”,而是采样预算会浪费在大量与当前 xx 无关的 zz 上。

真实后验却是:

pθ(z∣x)={1,z=17,0,z≠17.p_\theta(z\mid x)= \begin{cases} 1,&z=17,\\ 0,&z\ne17. \end{cases}

如果编码器已经学到近似后验 qϕ(z∣x)≈pθ(z∣x)q_\phi(z\mid x)\approx p_\theta(z\mid x),它会几乎总是采到 z=17z=17。于是重建项

Eqϕ(z∣x)[log⁡pθ(x∣z)]\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)]

可由极少样本稳定估计。

7.2 重参数化:让重建误差能够回传​

直接从依赖参数的分布采样,不便向编码器参数 ϕ\phi 反向传播。改为采样与参数无关的噪声:

ϵ∼N(0,I),z=μ+σ⊙ϵ=μ+exp⁡(12log⁡σ2)⊙ϵ.\epsilon\sim\mathcal N(0,I), \qquad z=\mu+\sigma\odot\epsilon =\mu+\exp\left(\tfrac12\log\sigma^2\right)\odot\epsilon.

给定 ϵ\epsilon 后,zz 是 μ,σ\mu,\sigma 的可微函数,梯度即可经由 zz 回传到编码器。

def reparameterize(mu, logvar):
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
return mu + eps * std

8. 最终训练损失:重建项与 KL 项​

对单张训练图片,最终要最小化的是负 ELBO:

LVAE(x)=−Eqϕ(z∣x)[log⁡pθ(x∣z)]⏟重建项+DKL(qϕ(z∣x)∥p(z))⏟KL 项.\mathcal L_{\mathrm{VAE}}(x) =\underbrace{-\mathbb E_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)]}_{\text{重建项}} +\underbrace{D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p(z)\bigr)}_{\text{KL 项}}.

实际训练时,对 batch 中每张 x(b)x^{(b)} 的 KK 个潜变量样本求平均,再对 BB 张图片求平均:

Lbatch=1B∑b=1B[−1K∑k=1Klog⁡pθ(x(b)∣z(b,k))+DKL(qϕ(z∣x(b))∥p(z))].\mathcal L_{\mathrm{batch}} =\frac1B\sum_{b=1}^{B} \left[ -\frac1K\sum_{k=1}^{K} \log p_\theta\bigl(x^{(b)}\mid z^{(b,k)}\bigr) +D_{\mathrm{KL}}\bigl(q_\phi(z\mid x^{(b)})\|p(z)\bigr) \right].

下面两个子节分别将这两项写成实际可计算的形式。

8.1 KL 项:高斯时可逐项解析计算​

重建项需要采样估计,但 KL 项在常用的高斯设定下不需要采样。对当前固定图片 xx,令潜变量维度为 dd:

qϕ(z∣x)=N(z;μ,diag⁡(σ2)),p(z)=N(z;0,I).q_\phi(z\mid x)=\mathcal N\bigl(z;\mu,\operatorname{diag}(\sigma^2)\bigr), \qquad p(z)=\mathcal N(z;0,I).

这里 μ,σ\mu,\sigma 都是编码器针对当前 xx 输出的 dd 维向量;μj,σj\mu_j,\sigma_j 表示第 jj 个维度。由于对角高斯的各维相互独立,KL 可以逐维相加。

从定义开始:

DKL(qϕ(z∣x)∥p(z))=∫qϕ(z∣x)[log⁡qϕ(z∣x)−log⁡p(z)]dz=Eqϕ[log⁡qϕ(z∣x)]−Eqϕ[log⁡p(z)].\begin{aligned} D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p(z)\bigr) &=\int q_\phi(z\mid x) \left[\log q_\phi(z\mid x)-\log p(z)\right]dz\\ &=\mathbb E_{q_\phi}[\log q_\phi(z\mid x)] -\mathbb E_{q_\phi}[\log p(z)]. \end{aligned}

先计算先验的对数概率。标准高斯满足:

log⁡p(z)=−d2log⁡(2π)−12∑j=1dzj2.\log p(z)=-\frac d2\log(2\pi)-\frac12\sum_{j=1}^{d}z_j^2.

而若 zj∼N(μj,σj2)z_j\sim\mathcal N(\mu_j,\sigma_j^2),则 Eq[zj2]=μj2+σj2\mathbb E_q[z_j^2]=\mu_j^2+\sigma_j^2。因此:

Eqϕ[log⁡p(z)]=−d2log⁡(2π)−12∑j=1dEq[zj2]=−d2log⁡(2π)−12∑j=1d(μj2+σj2).\begin{aligned} \mathbb E_{q_\phi}[\log p(z)] &=-\frac d2\log(2\pi) -\frac12\sum_{j=1}^{d}\mathbb E_q[z_j^2]\\ &=-\frac d2\log(2\pi) -\frac12\sum_{j=1}^{d}(\mu_j^2+\sigma_j^2). \end{aligned}
展开:标准高斯的对数概率与二阶矩从何而来

因为先验是 dd 维标准正态 p(z)=N(z;0,I)p(z)=\mathcal N(z;0,I),各维独立:

p(z)=∏j=1d12πexp⁡(−zj22)=(2π)−d/2exp⁡(−12∑j=1dzj2).\begin{aligned} p(z) &=\prod_{j=1}^{d}\frac{1}{\sqrt{2\pi}} \exp\left(-\frac{z_j^2}{2}\right)\\ &=(2\pi)^{-d/2}\exp\left(-\frac12\sum_{j=1}^{d}z_j^2\right). \end{aligned}

对上式取对数;乘积变为求和,指数项的对数就是指数:

log⁡p(z)=−d2log⁡(2π)−12∑j=1dzj2.\log p(z) =-\frac d2\log(2\pi)-\frac12\sum_{j=1}^{d}z_j^2.

另一方面,zj∼N(μj,σj2)z_j\sim\mathcal N(\mu_j,\sigma_j^2) 的均值和方差定义为:

Eq[zj]=μj,Var⁡q(zj)=Eq[(zj−μj)2]=σj2.\mathbb E_q[z_j]=\mu_j, \qquad \operatorname{Var}_q(z_j) =\mathbb E_q\left[(z_j-\mu_j)^2\right] =\sigma_j^2.

将平方展开:

σj2=Eq[zj2−2μjzj+μj2]=Eq[zj2]−2μjEq[zj]+μj2=Eq[zj2]−μj2.\begin{aligned} \sigma_j^2 &=\mathbb E_q[z_j^2-2\mu_jz_j+\mu_j^2]\\ &=\mathbb E_q[z_j^2]-2\mu_j\mathbb E_q[z_j]+\mu_j^2\\ &=\mathbb E_q[z_j^2]-\mu_j^2. \end{aligned}

移项即可得到:

Eq[zj2]=μj2+σj2.\boxed{\mathbb E_q[z_j^2]=\mu_j^2+\sigma_j^2}.

将 log⁡p(z)\log p(z) 代入期望,并利用期望的线性性,就得到正文中的 Eq[log⁡p(z)]\mathbb E_q[\log p(z)]。

再计算近似后验自身的对数概率。对角高斯满足:

log⁡qϕ(z∣x)=−d2log⁡(2π)−12∑j=1d[log⁡σj2+(zj−μj)2σj2].\log q_\phi(z\mid x) =-\frac d2\log(2\pi) -\frac12\sum_{j=1}^{d} \left[\log\sigma_j^2+\frac{(z_j-\mu_j)^2}{\sigma_j^2}\right].

由于 Eq[(zj−μj)2]=σj2\mathbb E_q[(z_j-\mu_j)^2]=\sigma_j^2,有:

Eqϕ[log⁡qϕ(z∣x)]=−d2log⁡(2π)−12∑j=1d(log⁡σj2+1).\begin{aligned} \mathbb E_{q_\phi}[\log q_\phi(z\mid x)] &=-\frac d2\log(2\pi) -\frac12\sum_{j=1}^{d}(\log\sigma_j^2+1). \end{aligned}

两式相减,−d2log⁡(2π)-\tfrac d2\log(2\pi) 抵消:

DKL(qϕ(z∣x)∥p(z))=12∑j=1d(μj2+σj2−1−log⁡σj2).\begin{aligned} D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p(z)\bigr) &=\frac12\sum_{j=1}^{d} \left(\mu_j^2+\sigma_j^2-1-\log\sigma_j^2\right). \end{aligned}

论文截图中写的是 ELBO 内的 −DKL-D_{\mathrm{KL}},故符号相反:

−DKL(qϕ(z∣x)∥p(z))=12∑j=1d(1+log⁡σj2−μj2−σj2).-D_{\mathrm{KL}}\bigl(q_\phi(z\mid x)\|p(z)\bigr) =\frac12\sum_{j=1}^{d} \left(1+\log\sigma_j^2-\mu_j^2-\sigma_j^2\right).

若编码器输出 logvar = \log\sigma^2,实现为:

kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())

8.2 重建项:固定方差高斯时对应 MSE​

若 pθ(x∣z)p_\theta(x\mid z) 是固定方差高斯,重建项的负对数似然(忽略常数与缩放)对应 MSE;若是 Bernoulli,则常用 BCE。下面推导前一种情形。

设一张图片被展平成 DD 维向量。解码器不直接输出一个确定的重建图,而是输出高斯分布的均值:

x^θ(z)=μθ(z),pθ(x∣z)=N(x;x^θ(z),σx2I).\hat{x}_\theta(z)=\mu_\theta(z), \qquad p_\theta(x\mid z)=\mathcal N\bigl(x;\hat{x}_\theta(z),\sigma_x^2I\bigr).

这不是说 xx 与 x^θ(z)\hat{x}_\theta(z) 相等,而是说:给定 zz 后,真实图片 xx 被模型视为从一个高斯分布中随机产生。其中:

  • N(变量;均值,协方差)\mathcal N(\text{变量};\text{均值},\text{协方差}) 是高斯分布的记号;分号左侧的 xx 是该分布要评价/生成的随机变量。
  • x^θ(z)\hat{x}_\theta(z) 是解码器输入 zz 后输出的预测图像,也是该高斯分布的均值。
  • σx2I\sigma_x^2I 是协方差矩阵:各像素独立,且都使用同一个方差 σx2\sigma_x^2。

等价的生成写法是:

x=x^θ(z)+η,η∼N(0,σx2I).x=\hat{x}_\theta(z)+\eta, \qquad \eta\sim\mathcal N(0,\sigma_x^2I).

这里观测方差 σx2\sigma_x^2 是预先固定的,不随 x,z,θx,z,\theta 改变。多维高斯密度为:

pθ(x∣z)=1(2πσx2)D/2exp⁡(−12σx2∥x−x^θ(z)∥22).p_\theta(x\mid z) =\frac{1}{(2\pi\sigma_x^2)^{D/2}} \exp\left( -\frac{1}{2\sigma_x^2} \left\|x-\hat{x}_\theta(z)\right\|_2^2 \right).

取负对数:

−log⁡pθ(x∣z)=D2log⁡(2πσx2)+12σx2∥x−x^θ(z)∥22.\begin{aligned} -\log p_\theta(x\mid z) &=\frac D2\log(2\pi\sigma_x^2) +\frac{1}{2\sigma_x^2} \left\|x-\hat{x}_\theta(z)\right\|_2^2. \end{aligned}

第一项与解码器参数 θ\theta 无关;第二项前的 12σx2>0\frac{1}{2\sigma_x^2}>0 也是固定常数。因此最小化负对数似然,与最小化:

∥x−x^θ(z)∥22=∑i=1D(xi−x^θ,i(z))2\left\|x-\hat{x}_\theta(z)\right\|_2^2 =\sum_{i=1}^{D}(x_i-\hat{x}_{\theta,i}(z))^2

有完全相同的最优解。若再除以 DD,就是通常代码里写的 MSE:

MSE⁡(x,x^)=1D∑i=1D(xi−x^i)2.\operatorname{MSE}(x,\hat{x}) =\frac1D\sum_{i=1}^{D}(x_i-\hat{x}_i)^2.

所以“MSE 重建损失”隐含了一个概率假设:给定 zz 后,每个像素等于解码器预测均值加上独立、同方差的高斯噪声。

实现时仍需明确重建项和 KL 项各自的 sum/mean 维度,否则二者的相对权重会随图片尺寸和 batch size 改变。

8.3 最终可实现的 batch loss​

令

x^(b,k)=x^θ(z(b,k)),logvar⁡j(b)=log⁡((σj(b))2).\hat{x}^{(b,k)}=\hat{x}_\theta\bigl(z^{(b,k)}\bigr), \qquad \operatorname{logvar}^{(b)}_j=\log\bigl((\sigma^{(b)}_j)^2\bigr).

将 batch 负 ELBO 中的重建项替换为 MSE、KL 项替换为对角高斯的闭式解后,忽略固定观测方差带来的常数与缩放,最终最小化的 loss 为:

Lbatch≈1B∑b=1B[1K∑k=1KMSE⁡(x(b),x^(b,k))+12∑j=1d((μj(b))2+exp⁡(logvar⁡j(b))−logvar⁡j(b)−1)].\boxed{ \begin{aligned} \mathcal L_{\mathrm{batch}} \approx\frac1B\sum_{b=1}^{B}\Bigg[ &\frac1K\sum_{k=1}^{K} \operatorname{MSE}\bigl(x^{(b)},\hat{x}^{(b,k)}\bigr)\\ &+\frac12\sum_{j=1}^{d} \left( (\mu^{(b)}_j)^2 +\exp\bigl(\operatorname{logvar}^{(b)}_j\bigr) -\operatorname{logvar}^{(b)}_j -1 \right) \Bigg]. \end{aligned} }

第一行是对每张图片的 KK 次潜变量采样取平均后的重建 MSE;第二行是该图片的正 KL 损失。最常见的设置是 K=1K=1:一个 batch 有 BB 张图片,每张图片经重参数化采一个 zz,算一次 MSE 和一次 KL,再对 BB 个样本取平均。

9. 训练与生成​

训练时,输入真实图片 xx:

  1. 编码器输出 μ(x),log⁡σ2(x)\mu(x),\log\sigma^2(x);
  2. 用重参数化得到 zz;
  3. 解码器根据 zz 输出 x^\hat{x} 或 pθ(x∣z)p_\theta(x\mid z) 的参数;
  4. 用重建项与 KL 项更新 θ,ϕ\theta,\phi。

生成时,不需要图片,也不使用编码器:

z∼p(z)=N(0,I)⟶x∼pθ(x∣z).z\sim p(z)=\mathcal N(0,I) \quad\longrightarrow\quad x\sim p_\theta(x\mid z).

KL 项的作用在这里闭环:训练时让 qϕ(z∣x)q_\phi(z\mid x) 接近 p(z)p(z),生成时从先验采到的 zz 才会位于解码器见过的潜空间区域。